3. Saving and Loading Models
Saving a model means storing its learned weights so they can be used later.
Loading a model restores those learned weights.
Save Model
The recommended approach is to save the model's state_dict.
torch.save(model.state_dict(), "model.pth")
Load Model
Create the same model architecture first, then load the saved weights.
model = NeuralNetwork()
model.load_state_dict(torch.load("model.pth"))
For inference:
model.eval()
State Dictionary
state_dict contains the model's learned parameters.
print(model.state_dict())
It mainly contains:
- Weights
- Biases
Save vs Load
| Save | Load |
|---|---|
torch.save() | torch.load() |
| Stores model parameters | Restores model parameters |
Usually saves state_dict | Loads into the same architecture |
Model Checkpoint
A checkpoint can store more than just model weights.
torch.save({
"model_state_dict": model.state_dict(),
"optimizer_state_dict": optimizer.state_dict(),
"epoch": epoch,
"loss": loss
}, "checkpoint.pth")
Useful when training needs to be resumed later.
Save Only Weights vs Checkpoint
Save state_dict | Save checkpoint |
|---|---|
| Mainly for inference/reuse | Mainly for continuing training |
| Smaller and simple | Stores training state too |
| Model architecture must be recreated | Model architecture and optimizer must be recreated |
Actual Implementation
import torch
import torch.nn as nn
import torch.optim as optim
X = torch.tensor([[1.0], [2.0], [3.0], [4.0]])
y = torch.tensor([[2.0], [4.0], [6.0], [8.0]])
class NeuralNetwork(nn.Module):
def __init__(self):
super().__init__()
self.linear = nn.Linear(1, 1)
def forward(self, x):
return self.linear(x)
model = NeuralNetwork()
optimizer = optim.SGD(model.parameters(), lr=0.01)
criterion = nn.MSELoss()
for epoch in range(1000):
output = model(X)
loss = criterion(output, y)
optimizer.zero_grad()
loss.backward()
optimizer.step()
# Save model
torch.save(model.state_dict(), "model.pth")
# Create a new model
loaded_model = NeuralNetwork()
# Load saved weights
loaded_model.load_state_dict(torch.load("model.pth", weights_only=True))
# Evaluation mode
loaded_model.eval()
with torch.no_grad():
predictions = loaded_model(X)
print(predictions)